feat(checkpoint_conversion): support direct FP8 and scale tensor ingestion in to_maxtext - #5052
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces support for FP8 weight-only storage and dynamic dequantization in MaxText, including changes to checkpoint conversion, linear layers, and decoder layers to support FP8 weights alongside scalar, per-channel, and block-wise scales. Feedback on these changes includes addressing a potential lossy double conversion of float16/bfloat16 tensors during checkpoint loading, extracting duplicated scale key resolution logic in to_maxtext.py into a helper function, and removing the redundant local _is_fp8_dtype helper in linears.py in favor of the existing utility in common_types.py.
| elif t.dtype == torch.bfloat16: | ||
| return t.to(torch.float32).numpy().astype(ml_dtypes.bfloat16) | ||
| elif t.dtype == torch.float16: | ||
| return t.to(torch.float32).numpy().astype(ml_dtypes.bfloat16) |
There was a problem hiding this comment.
Converting torch.float16 and torch.bfloat16 unconditionally to ml_dtypes.bfloat16 inside get_tensor can lead to a lossy double conversion when save_dtype is configured as float32 (i.e., float16 -> bfloat16 -> float32).
Since LazyTensor.__array__ already handles casting the retrieved numpy array to the target save_dtype (via arr.astype(dtype)), get_tensor should simply return the closest native numpy representation (e.g., float32 for bfloat16 to avoid PyTorch conversion issues, and native float16 for float16).
| elif t.dtype == torch.bfloat16: | |
| return t.to(torch.float32).numpy().astype(ml_dtypes.bfloat16) | |
| elif t.dtype == torch.float16: | |
| return t.to(torch.float32).numpy().astype(ml_dtypes.bfloat16) | |
| elif t.dtype == torch.bfloat16: | |
| return t.to(torch.float32).numpy() | |
| elif t.dtype == torch.float16: | |
| return t.numpy() |
| if shard_name is None: | ||
| # Check fallback for .weight_scale vs .scale and inverse scales | ||
| if resolved_key.endswith(".weight_scale"): | ||
| for suffix in [".scale", ".weight_scale_inv", ".scale_inv"]: | ||
| alt_key = resolved_key[:-len(".weight_scale")] + suffix | ||
| if alt_key in self.shard_map: | ||
| resolved_key = alt_key | ||
| shard_name = self.shard_map[resolved_key] | ||
| break | ||
| elif resolved_key.endswith(".scale"): | ||
| for suffix in [".weight_scale", ".scale_inv", ".weight_scale_inv"]: | ||
| alt_key = resolved_key[:-len(".scale")] + suffix | ||
| if alt_key in self.shard_map: | ||
| resolved_key = alt_key | ||
| shard_name = self.shard_map[resolved_key] | ||
| break |
There was a problem hiding this comment.
The fallback logic for resolving .weight_scale vs .scale and inverse scales is duplicated three times in this file (here, in get_tensor's second block, and in _eager_getter).
To improve maintainability and reduce redundancy, consider extracting this logic into a helper function at the module level:
def resolve_scale_key(key: str, container) -> str:
"""Resolves fallback keys for weight scales and inverse scales."""
if key in container:
return key
if key.endswith(".weight_scale"):
for suffix in [".scale", ".weight_scale_inv", ".scale_inv"]:
alt_key = key[:-len(".weight_scale")] + suffix
if alt_key in container:
return alt_key
elif key.endswith(".scale"):
for suffix in [".weight_scale", ".scale_inv", ".weight_scale_inv"]:
alt_key = key[:-len(".scale")] + suffix
if alt_key in container:
return alt_key
return keyYou can then simplify this block to:
resolved_key = resolve_scale_key(key, self.shard_map)
shard_name = self.shard_map.get(resolved_key)| def _is_fp8_dtype(dtype: Any) -> bool: | ||
| """Checks whether a dtype is an FP8 data type.""" | ||
| if dtype is None: | ||
| return False | ||
| try: | ||
| canon_dtype = _canonicalize_dtype(dtype) | ||
| except (TypeError, ValueError): | ||
| return False | ||
|
|
||
| fp8_types = [jnp.float8_e4m3fn, jnp.float8_e5m2] | ||
| for attr in ("float8_e4m3fnuz", "float8_e5m2fnuz", "float8_e4m3b11fnuz"): | ||
| if hasattr(jnp, attr): | ||
| fp8_types.append(getattr(jnp, attr)) | ||
| return canon_dtype in fp8_types |
There was a problem hiding this comment.
7687dfa to
f063fd7
Compare
56098db to
3d13d99
Compare
f063fd7 to
8e544e6
Compare
3d13d99 to
4e38337
Compare
8e544e6 to
f74726d
Compare
cb0ebb5 to
f54bfa9
Compare
|
|
||
| if "HF_HOME" not in os.environ and os.path.exists("/dev/shm"): | ||
| os.environ["HF_HOME"] = "/dev/shm/hf_cache" | ||
|
|
There was a problem hiding this comment.
remove this it was only for local development
9456383 to
c3e0942
Compare
f54bfa9 to
44cc3a7
Compare
| checkpointing.wait_until_finished(checkpoint_manager) | ||
|
|
||
| max_logging.log(f"Elapse for checkpoint save: {(time.time() - start) / 60:.2f} min") | ||
|
|
There was a problem hiding this comment.
delete this new line addition
| with safe_open(local_path, framework=framework, device="cpu") as f: | ||
| final_key = resolve_scale_key(resolved_key, f.keys()) | ||
| t = f.get_tensor(final_key) | ||
| if torch is not None and isinstance(t, torch.Tensor): | ||
| if hasattr(torch, "float8_e4m3fn") and t.dtype == torch.float8_e4m3fn: | ||
| return t.view(torch.uint8).numpy().view(ml_dtypes.float8_e4m3fn) | ||
| elif hasattr(torch, "float8_e5m2") and t.dtype == torch.float8_e5m2: | ||
| return t.view(torch.uint8).numpy().view(ml_dtypes.float8_e5m2) | ||
| elif t.dtype == torch.bfloat16: | ||
| if self.save_dtype in ("float32", DType.FLOAT32): | ||
| return t.to(torch.float32).numpy() | ||
| return t.view(torch.int16).numpy().view(ml_dtypes.bfloat16) | ||
| elif t.dtype == torch.float16: | ||
| if self.save_dtype in ("float32", DType.FLOAT32): | ||
| return t.to(torch.float32).numpy() | ||
| elif self.save_dtype in ("bfloat16", DType.BFLOAT16): | ||
| return t.to(torch.float32).numpy().astype(ml_dtypes.bfloat16) | ||
| return t.numpy() | ||
| else: | ||
| return t.numpy() | ||
| return t |
There was a problem hiding this comment.
can we simplify this somehow?
| if config.scan_layers: | ||
| # If it's a standard scanned layer, we use the configured param_scan_axis. | ||
| axis_to_stack = config.param_scan_axis | ||
| axis_to_stack = config.param_scan_axis if len(target_shape) > config.param_scan_axis else 0 |
There was a problem hiding this comment.
can you include a comment here?
My guess is this is used for 1D scales?
There was a problem hiding this comment.
Yes for 1D scales. Added a comment to specify this
2ae774c to
dce35b2
Compare
dbf4b02 to
fde4737
Compare
dce35b2 to
c8f282a
Compare
Codecov Report❌ Patch coverage is
📢 Thoughts on this report? Let us know! |
fde4737 to
ed14346
Compare
c8f282a to
82dc9d4
Compare
ed14346 to
205de7a
Compare
82dc9d4 to
b74c55f
Compare
205de7a to
41f7756
Compare
b74c55f to
279e76e
Compare
41f7756 to
ebc3f73
Compare
279e76e to
02eac58
Compare
ebc3f73 to
cb11b27
Compare
02eac58 to
6563931
Compare
cb11b27 to
5e61891
Compare
6563931 to
921921e
Compare
5e61891 to
f3a466f
Compare
921921e to
219da46
Compare
f3a466f to
e4f3ad5
Compare
219da46 to
a19e63c
Compare
e4f3ad5 to
69bc314
Compare
a19e63c to
d32cf49
Compare
69bc314 to
c84e000
Compare
d32cf49 to
1bb4531
Compare
…stion in to_maxtext Extends the MaxText HuggingFace checkpoint conversion pipeline to directly ingest FP8 weight tensors and companion scale tensors. Key Changes: - to_maxtext: Support float8_e4m3fn and float8_e5m2 save dtypes in _convert_tensor_to_numpy and _eager_getter, preserving each tensor's source precision (companion scales keep their source dtype; unquantized modules stay bfloat16). - to_maxtext: Support scale key suffix resolution (.weight_scale, .scale, .scale_inv, .weight_scale_inv) in LazyHFLoader and _eager_getter. - tensor_handling & to_maxtext: Handle 1D per-layer scalar scale stacking along axis 0 when scan_layers=True. - utils: Support slicing 1D MaxText parameters along axis 0 during checkpoint export. - tests: Add unit tests verifying FP8 tensor casting and scale suffix resolution in hf_checkpoint_conversion_test.py.
Description
Enhances MaxText's standalone checkpoint conversion tool (
src/maxtext/checkpoint_conversion/to_maxtext.py) to directly ingest, serialize, and store native 8-bit float weights (float8_e4m3fn) and companion scale tensors into Orbax format without upcasting.Motivation & Context
Previously,
to_maxtext.pyonly supportedsave_dtype="bfloat16"andsave_dtype="float32". Converting pre-quantized FP8 Hugging Face checkpoints forced CPU upcasting to BF16, doubling memory usage and storage footprint. Direct ingestion allows converting multi-gigabyte/terabyte FP8 checkpoints without memory bloat.Key Changes
torch.float8_e4m3fntensors asuint8views before bridging toml_dtypes.float8_e4m3fnin NumPy memory, bypassing standard NumPy float8 limitations.LazyHFLoaderto usesafetensorswithframework="pt"(when PyTorch is available) so native FP8 tensors can be deserialized from safetensors.ml_dtypes.float8_e4m3fnhandling toLazyTensordtype initialization.resolve_scale_key()to resolve.weight_scaleto.weight_scale_invfallbacks in bothLazyHFLoaderand_eager_getter, supporting checkpoints with inverse scale naming conventions (e.g. NeuralMagic / vLLM).--save_dtype="float8_e4m3fn",_convert_tensor_to_numpyonly converts FP8 weights toml_dtypes.float8_e4m3fn. Companion scale tensors and unquantized modules (e.g., embeddings, norm scales, router gates) retain their source precision (BF16 or FP32).convert_lora_to_maxtext_adapter, adapter weights remainbfloat16whensave_dtype="float8_e4m3fn".tensor_handling.py,utils.py,to_maxtext.py):_build_single_axis_stacked_tensor(tensor_handling.pyandto_maxtext.py), guarded stacking logic whenscan_layers=True: iflen(target_shape) <= param_scan_axis(such as 1D per-layer scalar scale tensors like Llama 3.1 FP8kernel_scale), tensors are stacked along axis 0 instead of attempting an out-of-bounds stack along axis 1 (param_scan_axis).process_maxtext_param(utils.py), added corresponding logic to slice along axis 0 for 1D parameters during MaxText parameter export._extract_conversion_args()and CLI argv preprocessing to allow conversion arguments (save_dtype,hf_model_path,lazy_load_tensors, etc.) to be supplied using standard MaxTextkey=valuesyntax without--prefixes, separating them cleanly from pyconfig arguments.--save_dtypechoices to includefloat8_e4m3fn.Part 2 of 5 in the FP8 Weight-Only Dynamic Dequantization series (depends on #5053).
Tests
Automated Unit Tests
Added unit test suite in
tests/unit/to_maxtext_test.py:ResolveScaleKeyTest: Verifies exact key matching,.weight_scale.weight_scale_invfallback, and unmatched key handling.ConvertTensorToNumpyTest: Verifiestorch.float8_e4m3fnml_dtypes.float8_e4m3fnzero-copy view, FP8ExtractConversionArgsTest: Verifies extraction and separation of conversion flags from mixed pyconfig CLI arguments.End-to-End Verification
Tested end-to-end checkpoint conversion of Hugging Face FP8 model (
neuralmagic/Meta-Llama-3.1-8B-Instruct-FP8):Verified that linear weights are stored in
float8_e4m3fnand scale tensors infloat32.Checklist
Before submitting this PR, please make sure (put X in square brackets):
gemini-reviewlabel.